Skip to content

Branch module

Branch

Bases: Module

Use a LanguageModel to select which module(s) to call based on an arbitrary input, a question and a list of labels.

The selected branch(es) output the data model computed using the inputs and module's branch, while the others output None. The output is always a tuple of length len(branches) so each label has a fixed positional slot regardless of which one was selected.

The behaviour of the selector depends on decision_type:

  • decision_type=Decision (default): exactly one branch is selected per call. All other slots are None.
  • decision_type=MultiDecision: one or more branches are selected per call. Non-selected slots remain None. Use this for multi-label routing where several branches may need to fire at once (e.g., an article that spans both science and finance, or a query that should be answered by both a retrieval and a tool-using sub-program).

Single-label example (one branch active per call):

import synalinks
import asyncio

async def main():
    class Query(synalinks.DataModel):
        query: str

    class Answer(synalinks.DataModel):
        answer: str

    class AnswerWithCritique(synalinks.DataModel):
        thinking: str
        critique: str
        answer: str

    language_model = synalinks.LanguageModel(
        model="ollama/mistral",
    )

    x0 = synalinks.Input(data_model=Query)
    (x1, x2) = await synalinks.Branch(
        question="What is the difficulty level of the above query?",
        labels=["easy", "difficult"],
        branches=[
            synalinks.Generator(
                data_model=Answer,
                language_model=language_model,
            ),
            synalinks.Generator(
                data_model=AnswerWithCritique,
                language_model=language_model,
            ),
        ],
        language_model=language_model,
    )(x0)
    x3 = x1 | x2

    program = synalinks.Program(
        inputs=x0,
        outputs=x3,
        name="adaptative_chain_of_thought",
        description="Useful to answer step by step only when needed",
    )

if __name__ == "__main__":
    asyncio.run(main())

Multi-label example (zero, one, or several branches active per call):

import synalinks
import asyncio

async def main():
    class Article(synalinks.DataModel):
        text: str

    class ScienceSummary(synalinks.DataModel):
        thinking: str
        science_summary: str

    class FinanceSummary(synalinks.DataModel):
        thinking: str
        finance_summary: str

    class SportsSummary(synalinks.DataModel):
        thinking: str
        sports_summary: str

    language_model = synalinks.LanguageModel(model="ollama/mistral")

    x0 = synalinks.Input(data_model=Article)
    # Each label has a fixed slot in the output tuple. With
    # MultiDecision, several may be populated at once; the rest
    # are None.
    (sci, fin, spo) = await synalinks.Branch(
        question="Which topics does this article cover?",
        labels=["science", "finance", "sports"],
        branches=[
            synalinks.Generator(
                data_model=ScienceSummary,
                language_model=language_model,
            ),
            synalinks.Generator(
                data_model=FinanceSummary,
                language_model=language_model,
            ),
            synalinks.Generator(
                data_model=SportsSummary,
                language_model=language_model,
            ),
        ],
        decision_type=synalinks.MultiDecision,
        language_model=language_model,
    )(x0)

if __name__ == "__main__":
    asyncio.run(main())

For a biotech-startup article the result might be (<ScienceSummary>, <FinanceSummary>, None): science and finance are both active, sports stays None. The non-active slots can be combined downstream with | (logical OR) the same way as in the single-label example.

Parameters:

Name Type Description Default
question str

The question to ask.

None
labels list

The list of labels to choose from (strings).

None
branches list

The list of modules or programs to select from.

None
inject_decision bool

If True, inject the decision to the branch inputs. (default to True).

True
return_decision bool

If True, return the decision with the branch outputs. (default to True).

True
language_model LanguageModel

The language model to use.

None
prompt_template str

The default jinja2 prompt template to use (see Generator).

None
examples list

The default examples to use in the prompt (see Decision).

None
instructions list

The default instructions to use (see Decision).

None
seed_instructions list

Optional. A list of instructions to use as seed for the optimization. If not provided, use the default instructions as seed.

None
temperature float

Optional. The temperature for the LM call.

None
max_tokens int

Optional. Default None (model's own default). Caps the generation length.

None
top_p float

Optional. Default None (model's own default). Nucleus sampling probability.

None
top_k int

Optional. Default None (model's own default). Top-k sampling cutoff.

None
reasoning_effort string

Optional. The reasoning effort for the LM call between ['minimal', 'low', 'medium', 'high', 'disable', 'none', None]. Default to None (no reasoning).

None
use_inputs_schema bool

Optional. Whether or not use the inputs schema in the decision prompt (Default to False) (see Decision).

False
use_outputs_schema bool

Optional. Whether or not use the outputs schema in the decision prompt (Default to False) (see Decision).

False
decision_type type

Optional. The decision module class. Defaults to Decision (single-label, exactly one branch active). Pass MultiDecision to enable multi-label routing where several branches may be active simultaneously.

Decision
name str

Optional. The name of the module.

None
description str

Optional. The description of the module.

None
trainable bool

Whether the module's variables should be trainable.

True
decision_model DecisionModel

Optional. A decision model to decide which branch to take, instead of the language model (see Decision). The branches keep their own models.

None
min_confidence float

Optional. With a decision model and a Decision, the confidence under which no branch is selected (see Decision).

None
threshold float

Optional. With a decision model and a MultiDecision, the probability from which a branch is selected (see MultiDecision).

None
Source code in synalinks/src/modules/core/branch.py
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
@synalinks_export(["synalinks.modules.Branch", "synalinks.Branch"])
class Branch(Module):
    """Use a `LanguageModel` to select which module(s) to call based on an
        arbitrary input, a question and a list of labels.

    The selected branch(es) output the data model computed using the inputs
    and module's branch, while the others output `None`. The output is
    always a tuple of length `len(branches)` so each label has a fixed
    positional slot regardless of which one was selected.

    The behaviour of the selector depends on `decision_type`:

    - `decision_type=Decision` (default): exactly **one** branch is
      selected per call. All other slots are `None`.
    - `decision_type=MultiDecision`: **one or more** branches are
      selected per call. Non-selected slots remain `None`. Use this for
      multi-label routing where several branches may need to fire at
      once (e.g., an article that spans both `science` and `finance`,
      or a query that should be answered by both a retrieval and a
      tool-using sub-program).

    Single-label example (one branch active per call):

    ```python
    import synalinks
    import asyncio

    async def main():
        class Query(synalinks.DataModel):
            query: str

        class Answer(synalinks.DataModel):
            answer: str

        class AnswerWithCritique(synalinks.DataModel):
            thinking: str
            critique: str
            answer: str

        language_model = synalinks.LanguageModel(
            model="ollama/mistral",
        )

        x0 = synalinks.Input(data_model=Query)
        (x1, x2) = await synalinks.Branch(
            question="What is the difficulty level of the above query?",
            labels=["easy", "difficult"],
            branches=[
                synalinks.Generator(
                    data_model=Answer,
                    language_model=language_model,
                ),
                synalinks.Generator(
                    data_model=AnswerWithCritique,
                    language_model=language_model,
                ),
            ],
            language_model=language_model,
        )(x0)
        x3 = x1 | x2

        program = synalinks.Program(
            inputs=x0,
            outputs=x3,
            name="adaptative_chain_of_thought",
            description="Useful to answer step by step only when needed",
        )

    if __name__ == "__main__":
        asyncio.run(main())
    ```

    Multi-label example (zero, one, or several branches active per call):

    ```python
    import synalinks
    import asyncio

    async def main():
        class Article(synalinks.DataModel):
            text: str

        class ScienceSummary(synalinks.DataModel):
            thinking: str
            science_summary: str

        class FinanceSummary(synalinks.DataModel):
            thinking: str
            finance_summary: str

        class SportsSummary(synalinks.DataModel):
            thinking: str
            sports_summary: str

        language_model = synalinks.LanguageModel(model="ollama/mistral")

        x0 = synalinks.Input(data_model=Article)
        # Each label has a fixed slot in the output tuple. With
        # MultiDecision, several may be populated at once; the rest
        # are None.
        (sci, fin, spo) = await synalinks.Branch(
            question="Which topics does this article cover?",
            labels=["science", "finance", "sports"],
            branches=[
                synalinks.Generator(
                    data_model=ScienceSummary,
                    language_model=language_model,
                ),
                synalinks.Generator(
                    data_model=FinanceSummary,
                    language_model=language_model,
                ),
                synalinks.Generator(
                    data_model=SportsSummary,
                    language_model=language_model,
                ),
            ],
            decision_type=synalinks.MultiDecision,
            language_model=language_model,
        )(x0)

    if __name__ == "__main__":
        asyncio.run(main())
    ```

    For a biotech-startup article the result might be
    `(<ScienceSummary>, <FinanceSummary>, None)`: `science` and
    `finance` are both active, `sports` stays `None`. The non-active
    slots can be combined downstream with `|` (logical OR) the same
    way as in the single-label example.

    Args:
        question (str): The question to ask.
        labels (list): The list of labels to choose from (strings).
        branches (list): The list of modules or programs to select from.
        inject_decision (bool): If True, inject the decision to the branch inputs.
            (default to True).
        return_decision (bool): If True, return the decision with the branch outputs.
            (default to True).
        language_model (LanguageModel): The language model to use.
        prompt_template (str): The default jinja2 prompt template
            to use (see `Generator`).
        examples (list): The default examples to use in the prompt
            (see `Decision`).
        instructions (list): The default instructions to use (see `Decision`).
        seed_instructions (list): Optional. A list of instructions to use as seed for the
            optimization. If not provided, use the default instructions as seed.
        temperature (float): Optional. The temperature for the LM call.
        max_tokens (int): Optional. Default None (model's own default). Caps the
            generation length.
        top_p (float): Optional. Default None (model's own default). Nucleus
            sampling probability.
        top_k (int): Optional. Default None (model's own default). Top-k sampling
            cutoff.
        reasoning_effort (string): Optional. The reasoning effort for the LM call
            between ['minimal', 'low', 'medium', 'high', 'disable', 'none', None].
            Default to None (no reasoning).
        use_inputs_schema (bool): Optional. Whether or not use the inputs
            schema in the decision prompt (Default to False) (see `Decision`).
        use_outputs_schema (bool): Optional. Whether or not use the outputs
            schema in the decision prompt (Default to False) (see `Decision`).
        decision_type (type): Optional. The decision module class. Defaults to
            `Decision` (single-label, exactly one branch active). Pass
            `MultiDecision` to enable multi-label routing where several
            branches may be active simultaneously.
        name (str): Optional. The name of the module.
        description (str): Optional. The description of the module.
        trainable (bool): Whether the module's variables should be trainable.
        decision_model (DecisionModel): Optional. A decision model to decide
            which branch to take, instead of the language model (see
            `Decision`). The branches keep their own models.
        min_confidence (float): Optional. With a decision model and a
            `Decision`, the confidence under which no branch is selected
            (see `Decision`).
        threshold (float): Optional. With a decision model and a
            `MultiDecision`, the probability from which a branch is selected
            (see `MultiDecision`).
    """

    def __init__(
        self,
        *,
        question=None,
        labels=None,
        branches=None,
        inject_decision=True,
        return_decision=True,
        language_model=None,
        prompt_template=None,
        examples=None,
        instructions=None,
        seed_instructions=None,
        temperature=None,
        max_tokens=None,
        top_p=None,
        top_k=None,
        reasoning_effort=None,
        use_inputs_schema=False,
        use_outputs_schema=False,
        decision_type=Decision,
        name=None,
        description=None,
        trainable=True,
        decision_model=None,
        min_confidence=None,
        threshold=None,
        **kwargs,
    ):
        super().__init__(
            name=name,
            description=description,
            trainable=trainable,
        )
        if not branches:
            raise ValueError("The `branches` argument must be provided.")
        if not isinstance(branches, list):
            raise ValueError("The `branches` must be a list of `Module` or `Program`.")
        if len(labels) != len(branches):
            raise ValueError("The `labels` and `branches` must have the same length.")
        self.question = question
        self.labels = labels
        self.branches = {labels[i]: m for i, m in enumerate(branches)}
        self.inject_decision = inject_decision
        self.return_decision = return_decision
        # Resolved here, from the arguments as given: the decision module
        # receives the resolved language model, so it could not tell the
        # default decision model apart from an explicit language model.
        resolved_decision_model = resolve_decision_model(decision_model, language_model)
        self.language_model = _get_lm(language_model)
        self.prompt_template = prompt_template
        self.examples = examples
        self.instructions = instructions
        self.seed_instructions = seed_instructions
        self.temperature = temperature
        self.max_tokens = max_tokens
        self.top_p = top_p
        self.top_k = top_k
        self.reasoning_effort = reasoning_effort
        self.use_inputs_schema = use_inputs_schema
        self.use_outputs_schema = use_outputs_schema
        # Only given when set, so a custom `decision_type` without decision
        # model support keeps working (the default decision model is only
        # given to a `decision_type` that takes one).
        decision_model_kwargs = {}
        if decision_model is not None or (
            resolved_decision_model is not None
            and "decision_model" in inspect.signature(decision_type).parameters
        ):
            decision_model_kwargs["decision_model"] = resolved_decision_model
        # Also only given when set: `min_confidence` is a `Decision` argument,
        # `threshold` a `MultiDecision` one.
        if min_confidence is not None:
            decision_model_kwargs["min_confidence"] = min_confidence
        if threshold is not None:
            decision_model_kwargs["threshold"] = threshold
        self.decision = decision_type(
            question=self.question,
            labels=self.labels,
            language_model=self.language_model,
            prompt_template=self.prompt_template,
            examples=self.examples,
            instructions=self.instructions,
            seed_instructions=self.seed_instructions,
            temperature=self.temperature,
            max_tokens=self.max_tokens,
            top_p=self.top_p,
            top_k=self.top_k,
            reasoning_effort=self.reasoning_effort,
            use_inputs_schema=self.use_inputs_schema,
            use_outputs_schema=self.use_outputs_schema,
            name="decision_" + self.name,
            **decision_model_kwargs,
        )
        self.decision_model = getattr(self.decision, "decision_model", None)
        self.min_confidence = min_confidence
        self.threshold = threshold

    async def call(self, inputs, training=False):
        outputs = [None] * len(self.branches)

        if not inputs:
            return tuple(outputs)

        decision = await self.decision(
            inputs,
            training=training,
        )

        if not decision:
            return tuple(outputs)

        choice = decision.get("choice", decision.get("choices"))

        if not choice:
            return tuple(outputs)

        if self.inject_decision:
            inputs = await ops.concat(
                inputs,
                decision,
                name="inputs_with_decision_" + self.name,
            )

        tasks = []

        async def execute_branch(
            inputs, module=None, decision=None, return_decision=False
        ):
            if not inputs:
                return None
            if return_decision:
                return await ops.logical_and(
                    decision,
                    await module(inputs),
                )
            else:
                return await module(inputs)

        for label in self.labels:
            module = self.branches[label]
            selected = False
            if isinstance(choice, str):
                if label == choice:
                    selected = True
            elif isinstance(choice, (list, set)):
                if label in choice:
                    selected = True
            if selected and module:
                tasks.append(
                    execute_branch(
                        inputs,
                        module,
                        decision,
                        return_decision=self.return_decision,
                    )
                )
            else:
                tasks.append(execute_branch(None))
        outputs = await asyncio.gather(*tasks)
        return tuple(outputs)

    async def compute_output_spec(self, inputs, training=False):
        outputs = []
        decision = await self.decision(
            inputs,
            training=training,
        )
        if self.inject_decision:
            inputs = await ops.concat(
                inputs,
                decision,
                name="inputs_with_decision_" + self.name,
            )
        for label in self.labels:
            module = self.branches[label]
            if self.return_decision:
                outputs.append(
                    await ops.logical_and(
                        decision,
                        await module(
                            inputs,
                            training=training,
                        ),
                        name="with_decision_" + self.name,
                    )
                )
            else:
                outputs.append(
                    await module(
                        inputs,
                        training=training,
                    )
                )
        return tuple(outputs)

    def get_config(self):
        config = {
            "question": self.question,
            "labels": self.labels,
            "inject_decision": self.inject_decision,
            "return_decision": self.return_decision,
            "prompt_template": self.prompt_template,
            "examples": self.examples,
            "instructions": self.instructions,
            "seed_instructions": self.seed_instructions,
            "temperature": self.temperature,
            "max_tokens": self.max_tokens,
            "top_p": self.top_p,
            "top_k": self.top_k,
            "reasoning_effort": self.reasoning_effort,
            "use_inputs_schema": self.use_inputs_schema,
            "use_outputs_schema": self.use_outputs_schema,
            "min_confidence": self.min_confidence,
            "threshold": self.threshold,
            "name": self.name,
            "description": self.description,
            "trainable": self.trainable,
        }
        language_model_config = {
            "language_model": serialization_lib.serialize_synalinks_object(
                self.language_model
            )
        }
        if self.decision_model is not None:
            language_model_config["decision_model"] = (
                serialization_lib.serialize_synalinks_object(self.decision_model)
            )
        branches_config = {
            "branches": [
                serialization_lib.serialize_synalinks_object(branch)
                for branch in self.branches.values()
            ]
        }
        return {**config, **language_model_config, **branches_config}

    @classmethod
    def from_config(cls, config, custom_objects=None):
        language_model = serialization_lib.deserialize_synalinks_object(
            config.pop("language_model")
        )
        if "decision_model" in config:
            config["decision_model"] = serialization_lib.deserialize_synalinks_object(
                config.pop("decision_model"),
            )
        branches = [
            serialization_lib.deserialize_synalinks_object(
                branch_config, custom_objects=custom_objects
            )
            for branch_config in config.pop("branches")
        ]
        return cls(language_model=language_model, branches=branches, **config)